#Fig.4
library(readxl)
library(dplyr)
library(ape)
library(phytools)
library(ggplot2)
library(ggrepel)

#AM 
rrdataAM <- as.data.frame(read_excel("Source Data_Fig. 4.xlsx", sheet = "Fig. 4a"))
traits_means <- rrdataAM %>%
  group_by(Species) %>%
  summarise(
    Mycorrhizal_types = first(Mycorrhizal_types),
    Exudation = mean(Exudation, na.rm = TRUE),
    RD = mean(RD, na.rm = TRUE),
    SRL = mean(SRL, na.rm = TRUE),
    RTD = mean(RTD, na.rm = TRUE),
    RNC = mean(RNC, na.rm = TRUE),
    .groups = "drop")

rrdataAM <- traits_means %>%
  dplyr::select(Species, Exudation, RD, SRL, RTD, RNC) %>%
  filter(if_all(c(Exudation, RD, SRL, RTD, RNC), ~ !is.na(.)))

rrdataAM <- as.data.frame(rrdataAM)
rownames(rrdataAM) <- rrdataAM$Species
rrdataAM[, c("Exudation","RD","SRL","RTD","RNC")] <- lapply(rrdataAM[, c("Exudation","RD","SRL","RTD","RNC")], as.numeric)

trait_mat <- as.matrix(rrdataAM[, c("Exudation","RD","SRL","RTD","RNC")])
tree <- read.tree("sp_treeAM.tre")

common_sp <- intersect(tree$tip.label, rownames(trait_mat))
tree_pruned <- drop.tip(tree, setdiff(tree$tip.label, common_sp))
trait_mat <- trait_mat[tree_pruned$tip.label, ]

ppca <- phyl.pca(tree_pruned, trait_mat, method = "lambda", mode = "corr")
scores_raw <- as.data.frame(ppca$S[, 1:2])
colnames(scores_raw) <- c("PC1", "PC2")
scores_raw$Species <- rownames(scores_raw)

loadings_raw <- as.data.frame(ppca$L[, 1:2])
colnames(loadings_raw) <- c("PC1", "PC2")
loadings_raw$variable <- rownames(loadings_raw)
rot_obj <- varimax(as.matrix(loadings_raw[, c("PC1", "PC2")]))

scores_rot <- as.matrix(scores_raw[, c("PC1", "PC2")]) %*% rot_obj$rotmat
load_rot   <- as.matrix(loadings_raw[, c("PC1", "PC2")]) %*% rot_obj$rotmat

theta <- -90 * pi / 180
Rmat <- matrix(c(cos(theta), -sin(theta), sin(theta),  cos(theta)),nrow = 2, byrow = TRUE)

scores_final   <- scores_rot %*% Rmat
loadings_final <- load_rot   %*% Rmat

scores_final[,1] <- -scores_final[,1]
loadings_final[,1] <- -loadings_final[,1]

scores_rot_df <- data.frame(PC1 = scores_final[,1], PC2 = scores_final[,2], Species = scores_raw$Species)
variable_rot <- data.frame(PC1 = loadings_final[,1] * 100, PC2 = loadings_final[,2] * 100, variable = loadings_raw$variable)

pAM <- ggplot(scores_rot_df, aes(PC1, PC2)) +
  geom_point(shape = 21, size = 4.5, stroke = 0.2, color = "black", fill = "#F39F4E") +
  labs(x = "Collaboration gradient", y = "Conservation gradient") +
  geom_segment(data = variable_rot, aes(x = 0, y = 0, xend = PC1, yend = PC2), arrow = arrow(length = unit(0.25, "cm"), type = "closed"), linewidth = 1, color = "black") +
  geom_text_repel(data = variable_rot %>% mutate(lab = ifelse(variable == "RNC", "Root [N]", variable)), aes(PC1, PC2, label = lab),
                  size = 6, color = "black", show.legend = FALSE) +
  scale_x_continuous(limits = c(-100, 100), expand = expansion(mult = 0.05)) +
  scale_y_continuous(limits = c(-100, 100), expand = expansion(mult = 0.05)) +
  annotate( "text", x = -100, y = 100, label = "AM", hjust = 0, vjust = 1, size = 7, fontface = "bold") +
  theme_classic(base_size = 18) +
  theme(
    axis.line = element_blank(),
    axis.ticks = element_line(linewidth = 0.8, color = "black"),
    axis.text = element_text(size = 16, color = "black"),
    axis.title = element_text(size = 18),
    panel.border = element_rect(linewidth = 0.9, color = "black", fill = NA),
    plot.margin = margin(8, 8, 8, 8)
  )

print(pAM)

#ECM
rrdataECM <- as.data.frame(read_excel("Source Data_Fig. 4.xlsx", sheet = "Fig. 4b"))
traits_means <- rrdataECM %>%
  group_by(Species) %>%
  summarise(
    Mycorrhizal_types = first(Mycorrhizal_types),
    Exudation = mean(Exudation, na.rm = TRUE),
    RD = mean(RD, na.rm = TRUE),
    SRL = mean(SRL, na.rm = TRUE),
    RTD = mean(RTD, na.rm = TRUE),
    RNC = mean(RNC, na.rm = TRUE),
    .groups = "drop")

rrdataECM <- traits_means %>%
  dplyr::select(Species, Exudation, RD, SRL, RTD, RNC) %>%
  filter(if_all(c(Exudation, RD, SRL, RTD, RNC), ~ !is.na(.)))

rrdataECM <- as.data.frame(rrdataECM)
rownames(rrdataECM) <- rrdataECM$Species
rrdataECM[, c("Exudation","RD","SRL","RTD","RNC")] <- lapply(rrdataECM[, c("Exudation","RD","SRL","RTD","RNC")], as.numeric)

trait_mat <- as.matrix(rrdataECM[, c("Exudation", "RD", "SRL", "RTD", "RNC")])
tree <- read.tree("sp_treeECM.tre")

common_sp <- intersect(tree$tip.label, rownames(trait_mat))
tree_pruned <- drop.tip(tree, setdiff(tree$tip.label, common_sp))
trait_mat <- trait_mat[tree_pruned$tip.label, ]

ppca <- phyl.pca(tree_pruned, trait_mat, method = "lambda", mode = "corr")
scores_raw <- as.data.frame(ppca$S[, 1:2])
colnames(scores_raw) <- c("PC1", "PC2")
scores_raw$Species <- rownames(scores_raw)

loadings_raw <- as.data.frame(ppca$L[, 1:2])
colnames(loadings_raw) <- c("PC1", "PC2")
loadings_raw$variable <- rownames(loadings_raw)
rot_obj <- varimax(as.matrix(loadings_raw[, c("PC1", "PC2")]))
load_rot <- as.matrix(loadings_raw[, c("PC1", "PC2")]) %*% rot_obj$rotmat

variable_rot <- as.data.frame(load_rot * 100)
colnames(variable_rot) <- c("PC1", "PC2")
variable_rot$variable <- loadings_raw$variable

scores_rot <- as.matrix(scores_raw[, c("PC1", "PC2")]) %*% rot_obj$rotmat
scores_rot_df <- as.data.frame(scores_rot)
colnames(scores_rot_df) <- c("PC1", "PC2")
scores_rot_df$Species <- scores_raw$Species

pECM <- ggplot(scores_rot_df, aes(PC1, PC2)) +
  geom_point(shape = 21, size = 4.5, stroke = 0.2, color = "black", fill = "#98CFE6") +
  labs(x = "Collaboration gradient", y = "Conservation gradient") +
  geom_segment(data = variable_rot, aes(x = 0, y = 0, xend = PC1, yend = PC2), arrow = arrow(length = unit(0.25, "cm"), type = "closed"), linewidth = 1, color = "black") +
  geom_text_repel(data = variable_rot %>% mutate(lab = dplyr::case_when( variable == "RNC" ~ "Root [N]", TRUE ~ variable)), aes(PC1, PC2, label = lab),
                  size = 6, fontface = "plain", color = "black", show.legend = FALSE) +
  scale_x_continuous(limits = c(-100, 100), expand = expansion(mult = 0.05)) +
  scale_y_continuous(limits = c(-100, 100), expand = expansion(mult = 0.05)) +
  annotate("text", x = -100, y = 100, label = "ECM", hjust = 0, vjust = 1, size = 7, fontface = "bold") +
  theme_classic(base_size = 18) +
  theme(
    axis.line = element_blank(),
    axis.ticks = element_line(linewidth = 0.8, color = "black"),
    axis.text = element_text(size = 16, color = "black"),
    axis.title = element_text(size = 18),
    panel.border = element_rect(linewidth = 0.9, color = "black", fill = NA),
    plot.margin = margin(8, 8, 8, 8)
  )

print(pECM)

Fig.4 <-ggarrange(pAM, pECM,
                  ncol =1, nrow = 2, 
                  labels = c("a", "b"), font.label=list(size=24),hjust = -0.2, vjust = 1) 
Fig.4 

ggsave("Fig.4.jpg", Fig.4, units = "cm", width = 13, height = 25, dpi = 1000)